import pandas as pd
import numpy as np
from scipy import stats
import os
import warnings
warnings.filterwarnings('ignore')

def load_and_preprocess_data(filepath):
    """Load and preprocess the main data file"""
    print("Loading data...")
    df = pd.read_csv(filepath, sep='\t')
    print(f"Original data shape: {df.shape}")
    
    # Basic preprocessing
    df.replace('.', np.nan, inplace=True)
    df.replace('NA', np.nan, inplace=True)
    df.replace('', np.nan, inplace=True)
    
    # Convert numeric columns
    numeric_cols = ['QUAL', 'GT_Freq', 'age', 'normalized_age', 'Quan-MMTV', 't_size']
    for col in numeric_cols:
        if col in df.columns:
            df[col] = pd.to_numeric(df[col], errors='coerce')
    
    # Print column names for debugging
    print("Available columns:")
    for col in df.columns:
        print(f"  - {col}")
    
    return df

def create_sample_level_data(df):
    """Create a sample-level dataset with one row per sample"""
    print("Creating sample-level data...")
    
    # Identify columns that are constant per sample
    sample_level_cols = [
        'FileID', 'FH', 'Her2', 'L/D', 'Lymph.Node', 'Progesteron-R', 
        'Estrogen-R', 'Qual-HPV', 'Qual-MMTV', 'Quan-MMTV', 'age', 'normalized_age',
        'GT_Freq_Categorized', 'Age_Categorized', 'FH_Processed', 'Estrogen-R_Processed',
        'Progesteron-R_Processed', 'Her2_Processed', 'Lymph.Node_Processed', 
        'Qual-MMTV_Processed', 'Qual-HPV_Processed', 'Quan-MMTV_Processed', 't_size',
        'bc_stage', 't_grade', 'bc_Microcalcifications', 't_size_processed',
        'bc_stage_processed', 't_grade_processed', 'bc_Microcal_processed'
    ]
    
    # Filter to only include columns that exist in the dataframe
    sample_level_cols = [col for col in sample_level_cols if col in df.columns]
    print(f"Using sample-level columns: {sample_level_cols}")
    
    # Get unique samples with their metadata
    sample_df = df[sample_level_cols].drop_duplicates(subset='FileID')
    
    return sample_df

def create_variant_level_data(df):
    """Create a variant-level dataset with statistics for each variant"""
    print("Creating variant-level data...")
    
    # Create binary indicator for variant presence (handle 1/2 genotype)
    df['Variant_Present'] = df['GT'].apply(lambda x: 1 if str(x) in ['0/1', '1/1', '1/2'] else 0)
    
    # Define aggregation functions
    agg_functions = {
        'Variant_Present': ['count', 'sum'],
        'FH': 'first',
        'GT_Freq': 'mean',
        'QUAL': 'mean'
    }
    
    # Add MAF_gnomAD if it exists
    if 'MAF_gnomAD' in df.columns:
        agg_functions['MAF_gnomAD'] = 'first'
    
    # Group by variant and calculate statistics
    groupby_cols = ['CHROM', 'POS', 'REF', 'ALT', 'POS_ID', 'dbSNP/Novel', 'GENE']
    groupby_cols = [col for col in groupby_cols if col in df.columns]
    
    variant_stats = df.groupby(groupby_cols).agg(agg_functions).reset_index()
    
    # Flatten column names
    variant_stats.columns = ['_'.join(col).strip('_') if col[1] else col[0] for col in variant_stats.columns]
    
    # Rename columns
    rename_dict = {
        'Variant_Present_count': 'Total_Samples',
        'Variant_Present_sum': 'Samples_With_Variant',
        'FH_first': 'FH_Info',
        'GT_Freq_mean': 'Mean_GT_Freq',
        'QUAL_mean': 'Mean_QUAL'
    }
    
    if 'MAF_gnomAD_first' in variant_stats.columns:
        rename_dict['MAF_gnomAD_first'] = 'MAF_gnomAD'
    
    variant_stats.rename(columns=rename_dict, inplace=True)
    
    # Calculate variant frequency
    variant_stats['Variant_Frequency'] = variant_stats['Samples_With_Variant'] / variant_stats['Total_Samples']
    
    return variant_stats

def create_fh_counts(df):
    """Create counts of variants in FH positive vs FH negative samples"""
    print("Creating FH counts...")
    
    # Create binary indicator for variant presence (handle 1/2 genotype)
    df['Variant_Present'] = df['GT'].apply(lambda x: 1 if str(x) in ['0/1', '1/1', '1/2'] else 0)
    
    # Create FH indicator (assuming 1=FH positive, 0=FH negative)
    df['Is_FH'] = df['FH'].apply(lambda x: 1 if x == 1 else 0)
    
    # Group by variant and FH status
    fh_counts = df.groupby(['POS_ID', 'Is_FH']).agg({
        'Variant_Present': ['count', 'sum']
    }).reset_index()
    
    # Flatten column names
    fh_counts.columns = ['_'.join(col).strip('_') if col[1] else col[0] for col in fh_counts.columns]
    fh_counts.rename(columns={
        'Variant_Present_count': 'Total_Samples',
        'Variant_Present_sum': 'Samples_With_Variant'
    }, inplace=True)
    
    # Pivot to have FH positive and negative as separate columns
    fh_pivot = fh_counts.pivot_table(
        index='POS_ID',
        columns='Is_FH',
        values=['Total_Samples', 'Samples_With_Variant'],
        fill_value=0
    ).reset_index()
    
    # Flatten column names
    fh_pivot.columns = [
        'POS_ID', 
        'FH_Negative_Total', 'FH_Positive_Total', 
        'FH_Negative_With_Variant', 'FH_Positive_With_Variant'
    ]
    
    # Calculate odds ratios and confidence intervals
    fh_pivot['Odds_Ratio'] = fh_pivot.apply(
        lambda row: calculate_odds_ratio(
            row['FH_Positive_With_Variant'], 
            row['FH_Positive_Total'] - row['FH_Positive_With_Variant'],
            row['FH_Negative_With_Variant'], 
            row['FH_Negative_Total'] - row['FH_Negative_With_Variant']
        )[0], axis=1
    )
    
    fh_pivot['CI_Lower'] = fh_pivot.apply(
        lambda row: calculate_odds_ratio(
            row['FH_Positive_With_Variant'], 
            row['FH_Positive_Total'] - row['FH_Positive_With_Variant'],
            row['FH_Negative_With_Variant'], 
            row['FH_Negative_Total'] - row['FH_Negative_With_Variant']
        )[1][0], axis=1
    )
    
    fh_pivot['CI_Upper'] = fh_pivot.apply(
        lambda row: calculate_odds_ratio(
            row['FH_Positive_With_Variant'], 
            row['FH_Positive_Total'] - row['FH_Positive_With_Variant'],
            row['FH_Negative_With_Variant'], 
            row['FH_Negative_Total'] - row['FH_Negative_With_Variant']
        )[1][1], axis=1
    )
    
    # Merge with variant information
    variant_info_cols = ['POS_ID', 'CHROM', 'GENE', 'REF', 'ALT', 'dbSNP/Novel']
    if 'MAF_gnomAD' in df.columns:
        variant_info_cols.append('MAF_gnomAD')
        
    variant_info = df[variant_info_cols].drop_duplicates()
    fh_pivot = fh_pivot.merge(variant_info, on='POS_ID', how='left')
    
    return fh_pivot

def calculate_odds_ratio(a, b, c, d):
    """Calculate odds ratio and 95% confidence interval"""
    # Add 0.5 to all cells to avoid division by zero
    a, b, c, d = a+0.5, b+0.5, c+0.5, d+0.5
    
    # Calculate odds ratio
    odds_ratio = (a * d) / (b * c)
    
    # Calculate standard error
    se = np.sqrt(1/a + 1/b + 1/c + 1/d)
    
    # Calculate confidence interval
    ci_lower = np.exp(np.log(odds_ratio) - 1.96 * se)
    ci_upper = np.exp(np.log(odds_ratio) + 1.96 * se)
    
    return odds_ratio, (ci_lower, ci_upper)

def create_genotype_matrix(df):
    """Create a genotype matrix with samples as rows and variants as columns"""
    print("Creating genotype matrix...")
    
    # Create a binary indicator for variant presence (handle 1/2 genotype)
    df['Genotype_Code'] = df['GT'].apply(lambda x: 1 if str(x) in ['0/1', '1/1', '1/2'] else 0)
    
    # Pivot to create genotype matrix
    genotype_matrix = df.pivot_table(
        index='FileID',
        columns='POS_ID',
        values='Genotype_Code',
        fill_value=0,
        aggfunc='max'  # Use max to ensure we get 1 if variant is present in any of the duplicate rows
    ).reset_index()
    
    return genotype_matrix

def main():
    # File path
    filepath = "/home/mahmo/TUMOR/latest/forest/last/maindata.tsv"
    
    # Output directory
    output_dir = "/home/mahmo/TUMOR/latest/forest/last/transformed_data"
    
    # Create output directory if it doesn't exist
    os.makedirs(output_dir, exist_ok=True)
    
    # Load and preprocess data
    df = load_and_preprocess_data(filepath)
    
    # Create sample-level data
    sample_df = create_sample_level_data(df)
    
    # Create variant-level data
    variant_stats = create_variant_level_data(df)
    
    # Create FH counts
    fh_df = create_fh_counts(df)
    
    # Create genotype matrix
    genotype_matrix = create_genotype_matrix(df)
    
    # Merge sample data with genotype matrix
    final_sample_data = sample_df.merge(genotype_matrix, on='FileID', how='left')
    
    # Save all datasets
    print("Saving transformed data...")
    sample_df.to_csv(os.path.join(output_dir, "sample_level_data.tsv"), sep='\t', index=False)
    variant_stats.to_csv(os.path.join(output_dir, "variant_level_data.tsv"), sep='\t', index=False)
    fh_df.to_csv(os.path.join(output_dir, "fh_data.tsv"), sep='\t', index=False)
    final_sample_data.to_csv(os.path.join(output_dir, "final_sample_data.tsv"), sep='\t', index=False)
    
    print("Transformation complete!")
    print(f"Sample-level data shape: {sample_df.shape}")
    print(f"Variant-level data shape: {variant_stats.shape}")
    print(f"FH data shape: {fh_df.shape}")
    print(f"Final sample data shape: {final_sample_data.shape}")
    print(f"Files saved to: {output_dir}")

if __name__ == "__main__":
    main()